Add per-sequence causal policy to packed THD attention - #3274
Conversation
| @pytest.mark.parametrize( | ||
| "sequence_is_causal", | ||
| ( | ||
| (False, True, False, True), |
There was a problem hiding this comment.
The PyTorch QA job explicitly enumerates attention test modules and does not include this new file, so its forward, gradient, validation, and empty-batch coverage never runs in CI and regressions in the new dispatch path go undetected.
Knowledge Base Used: Tests and QA
Greptile SummaryThe PR adds per-sequence mask-policy dispatch for packed THD attention and registers its focused tests in the PyTorch L0 CI suite.
Confidence Score: 5/5The PR appears safe to merge. No blocking failure remains; the previously reported CI coverage omission is fixed by the new L0 pytest invocation. Important Files Changed
Flowchart%%{init: {'theme': 'neutral'}}%%
flowchart LR
A["Packed THD Q/K/V"] --> B{"Policy mapping"}
B --> C["Scalar-policy attention call: full context"]
B --> D["Scalar-policy attention call: causal"]
C --> E["Mask unselected token outputs"]
D --> F["Mask unselected token outputs"]
E --> G["Add disjoint outputs"]
F --> G
G --> H["Original packed token order"]
Reviews (4): Last reviewed commit: "Mask inactive mixed-THD lanes for FA3" | Re-trigger Greptile |
| group_max_seqlen, | ||
| is_causal=is_causal, | ||
| ) | ||
| output = output.index_copy(0, group_token_indices, group_output) |
There was a problem hiding this comment.
The miscellaneous kernels generated in this approach might degrade the overall performance quite a bit. TE supports an option called pad_between_seqs=True, which allows attention backends to skip through the padded tokens and only process the "intended" tokens/sequences. In your case, I think we can perform two attention calls, one viewing the causal sequences as the "intended" sequences (and the rest "pad tokens"), and the other one viewing the non-causal sequences as the "intended" sequences. Could you please take a look at this note and this test see if it makes sense for your use case?
For example, we have a batch of 7 sequences, [111aaa222bb3344cccc], where we would like to perform "padding_causal" on the numbered sequences, {111, 222, 33, 44}, and "padding" on the lettered ones, {aaa, bb, cccc}, we can do something like this:
mask_type='padding_causal'
cu_seqlens=[0,3,6,8,10]
cu_seqlens_padded=[0,6,11,13,19]
out_causal = self.forward(
q, k, v, # full batch, not split
qkv_format='thd',
attn_mask_type=mask_type,
cu_seqlens_q=cu_seqlens,
cu_seqlens_kv=cu_seqlens,
cu_seqlens_q_padded=cu_seqlens_padded,
cu_seqlens_kv_padded=cu_seqlens_padded,
) # shape [19, h, d] with 0s on pad positions
mask_type='padding'
cu_seqlens=[0,3,5,9]
cu_seqlens_padded=[3,9,15,19]
out_non_causal = self.forward(
q, k, v, # full batch, not split
qkv_format='thd',
attn_mask_type=mask_type,
cu_seqlens_q=cu_seqlens,
cu_seqlens_kv=cu_seqlens,
cu_seqlens_q_padded=cu_seqlens_padded,
cu_seqlens_kv_padded=cu_seqlens_padded,
) # shape [19, h, d] with 0s on pad positions
out = out_causal + out_non_causal
| restores the original packed token order. It therefore shares the | ||
| surrounding encoder computation but is not a single fused-attention | ||
| kernel. Initial support is limited to FP16/BF16 THD self-attention | ||
| without attention dropout, context parallelism, cacheing, FP8, |
There was a problem hiding this comment.
Nit: cacheing -> caching
But more importantly, can we rename the argument to something more generic and extensible, for example, attn_mask_type_per_seq? We can have the type as a dict, for example, users can pass in {"padding": torch.Tensor, "padding_causal": xxx, "padding_causal_bottom_right": xxx, xxx} to specify which SeqIDs need to run with "padding", and which with "padding_causal", etc. The tensor needs to be on-device.
| raise ValueError( | ||
| "thd_sequence_is_causal requires THD Q, K, and V with the same token count." | ||
| ) | ||
| if int(cu_seqlens_q[-1].item()) != query_layer.shape[0]: |
There was a problem hiding this comment.
Please try to avoid GPU-CPU syncs (e.g. item()). I think if we pursue the approach suggested above, we can achieve CUDA graph and torch.compile compatibility.
| if effective_bottom_right_diagonal is True: | ||
| raise ValueError( | ||
| "thd_sequence_is_causal requires the standard top-left causal diagonal." | ||
| ) |
There was a problem hiding this comment.
Some of the runtime checks can stay here, but others, regarding what's supported and what's not, can probably go to get_attention_backend. For example, the feature only supports THD, non-CP, etc. If you pursue the out=out_causal+out_non_causal approach above, the support matrix will expand naturally I think, for example, with self/cross-attention, bottom_right_diagonal=T/F. Also, pad_between_seqs=T already has some constraints in get_attention_backend which can be reused for this feature.
| else: | ||
| seqlens_kv = cu_seqlens_kv[1:] - cu_seqlens_kv[:-1] | ||
| max_seqlen_kv = int((seqlens_kv.max().item() + 63) // 64 * 64) | ||
|
|
There was a problem hiding this comment.
The entire if thd_sequence_is_causal is not None code block can sit here, to take advantage of the existing mask/window_size checks, as well as be in the self.prepare_forward_ctx context for FP8 (if FP8 is ever available for this feature).
Signed-off-by: Desh Raj <r.desh26@gmail.com>
Signed-off-by: Desh Raj <r.desh26@gmail.com>
Replace the per-policy Q/K/V gather and output scatter dispatcher with the review-suggested inter-sequence-padding design. Keep the original packed Q/K/V storage, derive logical and physical cu-seqlens for each mask policy, run the ordinary scalar-mask DPA path with pad_between_seqs, and sum its disjoint zero-padded outputs. Expose a generic attn_mask_type_per_seq mapping plus window_size_per_mask_type so callers can combine full-context offline attention with bounded causal attention. Preserve the scalar fast path for uniform batches, avoid policy-dependent host synchronization, and cover forward/backward parity, existing physical padding, cross-attention, per-policy windows, bottom-right masks, invalid inputs, empty batches, and CUDA graph replay. NRT H100 validation passed all 16 TransformerEngine cases and 3 focused Megatron integration tests. In an 8-node/64-H100 matched-auto comparison over iterations 7-100, grouped averaged 1802.289 ms and pad-between averaged 1779.596 ms (-1.259%); the paired 95% interval crossed zero, indicating practical parity when the backend is held constant. In the deployable-backend comparison, grouped/Flash averaged 1596.129 ms while pad-between/auto-cuDNN averaged 1766.769 ms: pad-between was 10.691% slower by mean and 11.482% slower by 5%-trimmed mean. The paired penalty was +170.640 ms with a 95% interval of +89.728 to +251.552 ms, with all 16 fully logged workload counters matching across 100 steps. The implementation removes the original data-movement contention, but the current image cannot run full-context inter-sequence padding through FlashAttention 2 and has no FA3, so the path falls back to cuDNN fused attention. Retain grouped/Flash for production until a competitive padding-capable backend is available; this commit keeps the cleaner generic API for review and future backend support. Signed-off-by: Desh Raj <r.desh26@gmail.com>
faea818 to
a88c9b6
Compare
for more information, see https://pre-commit.ci
Mixed-mask dispatch represents sequences owned by other policies as inter-sequence padding. FA3 preserves that layout contract, but its padding lanes are unspecified: focused Hopper testing observed NaNs in inactive forward outputs and tile-spill garbage in inactive dQ/dK/dV lanes. Build a sync-free mask from the physical cu_seqlens, sanitize each policy output with torch.where, and wrap Q/K/V in an identity-forward autograd guard that zeros only inactive gradient lanes. This keeps the pad_between_seqs design, avoids regrouping or copying tokens, and lets disjoint policy results be summed safely across FA2 and FA3. Validation: NRT H100 source-overlay run passed all 16 focused mixed-THD tests for FP16/BF16 forward and backward, pre-existing padding, cross attention, per-policy windows, CUDA graphs, and validation errors. Signed-off-by: Desh Raj <r.desh26@gmail.com>
Description
Add an optional per-sequence causal policy to
DotProductAttentionfor packed THD self-attention. A scalarattn_mask_typecurrently applies one policy to the entire packed invocation, while mixed offline/causal training needs each packed sequence to retain its own policy.When
thd_sequence_is_causalis supplied, Transformer Engine groups whole sequences by policy, invokes the existing scalar-policy attention path once per non-empty group, and restores the original packed token order. Uniform all-causal and all-offline batches retain the existing one-call path.This is a logical dispatcher over existing attention calls, not a new heterogeneous-mask fused kernel. Initial support is intentionally limited to plain FP16/BF16 THD self-attention; options whose tensors or semantics cannot be safely regrouped are rejected explicitly.
Type of change
Changes
thd_sequence_is_causal, a boolean tensor with one value per packed sequence, toDotProductAttention.forward.Validation
python -m pytest --import-mode=importlib tests/pytorch/attention/test_mixed_thd_attention.py -q -vvattention
python -m py_compile, andgit diff --checkChecklist:
The focused GPU suite and repository Python lint pass as described above; the full repository unit-test suite was not run locally.